Repository navigation
[WS1][CUDA][MiniMax-H3] ws1_one_h3_block: one real block forward/backward with first-drift report (RFC #420 stack 10/10) - #539
Open
fusheng-ji wants to merge 148 commits into
Open
fusheng-ji wants to merge 148 commits into
fusheng-ji wants to merge 148 commits into
Conversation
…+ shared H3 harness RFC RL-Align#420 step 3, row `timestep_sinusoid_h3`. - CUDA kernel: one thread per (t, k), the provider's exact FP32 op sequence (__fmul_rn, precise expf/sinf/cosf); bitwise equal to diffusers' get_timestep_embedding on CUDA. Analytic row-local backward. Non-1-D, empty, integer, non-finite or out-of-[0, 1] timesteps fail closed (RFC probe H10). - PyTorch reference (provider replay + FP64 golden), registry, gtest. - Shared H3 harness used by the following rows: * h3_manifest.json + prepare_h3_weights.py: pinned revision and sha256, only shard 1 of 14 is downloaded; * h3_provider.py: diffusers replayed op for op; * h3_chain.py: stage-wise chain replay, each stage declaring its promises (provider-bitwise, golden tolerance); used by scripts/h3_chain_replay.py and tests/h3/test_h3_conditioning_e2e.py; * h3_report.py + scripts/h3_evidence.py + scripts/plot_h3_evidence.py: per-op performance/accuracy JSON and figure; * benchmarks/benchmark_h3_conditioning.py. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Produced by scripts/h3_evidence.py at 0522865 on a clean tree; rendered with scripts/plot_h3_evidence.py and embedded in the operator doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…p MLP RFC RL-Align#420 step 3, row `timestep_mlp_fp32` (stacked on timestep_sinusoid_h3). - New CUDA module csrc/cuda/h3/det_linear.cu, contract h3-det-linear-v1: warp per output-column group, lane-strided 16-byte K chunks, ascending fmaf chains, xor butterfly, bias then fused SiLU. The launch shape never changes a column's summation order, so rows are batch/position invariant and runs are repeat-bitwise. - Backward without cuBLAS or atomics: dinput via fixed 64-row N chunks folded in order (row-local), dweight/dbias as ascending-row FP32 folds. - PyTorch reference (provider path + FP64 golden); BF16 anywhere is rejected (RFC probe H7: the time_embedder is a declared FP32 module). - On the pinned weights both CUDA and cuBLAS are ~1e-6 from the FP64 golden (100x inside the contract); the op is 2x faster for T >= 2. - Chain stage (covered by the end-to-end test), report/figure hooks, registry, gtest, doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Produced by scripts/h3_evidence.py at 65ef7f6 on a clean tree; rendered with scripts/plot_h3_evidence.py and embedded in the operator doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…rojection RFC RL-Align#420 step 3, row `adaln_projection_3mod` (stacked on timestep_mlp_fp32). - SiLU at temb's FP32 precision and exactly one cast to the projection dtype (bitwise equal to the provider activation); a BF16 temb is rejected (RFC probe H7). - 2688 -> 96768 projection on a BF16 tensor-core path in det_linear.cu, contract h3-det-linear-bf16-mma-v1: warp = 16 output columns x 8 input rows (zero-padded), K in ascending 16-wide groups, each mma.sync m16n8k16 from a zero accumulator added into an FP32 running sum. Same instruction sequence for every T <= 8, so rows are bitwise batch/order/tile invariant; 99.98% of outputs correctly rounded vs 99.84% for cuBLAS; kernel time flat at ~78 us for T = 1..8. SM80+. - Six (3T, 5376) views in diffusers' view(-1, 6H).chunk(6) layout, row t*3+m; a channel-tagging test pins chunk/modality placement (probe H4). - Backward: FP32 d_temb through the identity-VJP cast, ascending-row dW/db folds, no cuBLAS or atomics. The FP64 golden applies the cast straight-through so its own gradient is not rounded to BF16. - Chain stage (end-to-end test), report/figure hooks, registry, gtest, doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Produced by scripts/h3_evidence.py at d06e120 on a clean tree; rendered with scripts/plot_h3_evidence.py and embedded in the operator doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
RFC RL-Align#420 step 3, row `adaln_row_gather` (stacked on adaln_projection_3mod). - Forward: one launch gathers all six (S, H) modulation tensors by timestep_index * 3 + token_tag with 16-byte copies; bitwise equal to diffusers' six index_select calls. Out-of-range tags or timesteps fail closed (RFC probes H2/H3), never clamped. - Backward: deterministic FP32 segmented sum (stable sort by row, fixed 256-position tiles, ordered fold, one cast; no atomics) instead of index_select's atomic BF16 scatter-add. The PyTorch reference gets an equally deterministic backward. - H3AdaLNModulationCudaOp: projection + gather as one autograd node, bitwise-equal forward, FP32 table gradient into the projection backward. - Chain replay gains the gather stage and a backward mode (separate, fused and diffusers chains vs an FP64 golden); the end-to-end test now covers the whole chain forward and backward. - Report/figure hooks, registry, gtest (reduction class: the VJP is a segmented sum), doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…y and figure Produced by scripts/h3_evidence.py and scripts/h3_chain_replay.py --backward at fa551c2 on a clean tree; rendered with scripts/plot_h3_evidence.py and embedded in the operator doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
RFC RL-Align#420 step 4, row `h3_rmsnorm` (stacked on adaln_row_gather). - CUDA kernel replaying PyTorch's own RMSNorm reduction order (vectorized_layer_norm_kernel: (32, 4) block per row, 4-element vectors, shuffle-down and cross-warp trees, rsqrtf, w * (rstd * x)): bitwise equal to nn.RMSNorm on all four pinned H3 norm weights and batch/position invariant. - Fused AdaLN modulation n * (1 + scale[i]) + shift[i] with the rows gathered in-kernel from strided table views and each step rounded where the eager expression rounds: bitwise equal to diffusers for the block (adaln_indices) and norm_out (timestep_indices), without materialising the (S, H) gathered tensors; 2.5x faster forward at S = 131072. - Deterministic backward: row-local dx, dweight over fixed 256-row tiles, dshift/dscale as sorted segment sums (diffusers: non-deterministic BF16 atomics, ~15x larger error). - Manifest pins the block/refiner-final/norm_out norm weights and norm_out.linear (all in shard 1); chain gains a per-mode history and the block norm1 + modulation stage (covered by the end-to-end test). - Registry, gtest, report/figure hooks, doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Produced by scripts/h3_evidence.py at f36aa60 on a clean tree; rendered with scripts/plot_h3_evidence.py and embedded in the operator doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…ernel gate gather RFC RL-Align#420 step 6, row `adaln_gate_residual` (stacked on h3_rmsnorm). - residual + gate[adaln_index] * sublayer_output in one pass, the gate row read from the AdaLN table view in the kernel; the two roundings sit where the eager expression rounds (__fmul_rn/__fadd_rn, no FMA contraction), so the forward is bitwise equal to diffusers in bf16/fp16/fp32. 1.7x faster forward at S = 131072. - Backward: d_residual = grad; d_sublayer bitwise equal to the eager VJP (exact product, one rounding); d_gate a deterministic FP32 segment sum (diffusers: non-deterministic BF16 atomics, 13-49x larger error, varying run to run). The PyTorch reference gets an equally deterministic gate gradient. - Chain stage residual + gate_msa * (stand-in attention output), covered by the end-to-end test; registry, gtest, report/figure hooks, doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Produced by scripts/h3_evidence.py at a26b41a on a clean tree; rendered with scripts/plot_h3_evidence.py and embedded in the operator doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…norm/modulation RFC RL-Align#420 step 7, row `final_adaln_out` (stacked on adaln_gate_residual). - norm_out as one autograd node: FP32 SiLU -> one BF16 cast -> the deterministic tensor-core GEMV on norm_out.linear (shift first), then the h3_rmsnorm kernel indexed by timestep_indices. Forward equal to diffusers on 99.96% of elements (1-ULP projection ties); rows bitwise batch/position invariant; 2.4x faster forward at S = 131072. - Backward keeps the table gradient in FP32 into the projection backward: d_temb/dW/db 13-57x closer to FP64 than diffusers (whose error varies run to run), all gradients repeat-bitwise. - Golden rounds only at the declared boundaries (SiLU cast, BF16 table), straight-through for the gradient; gtest inputs at realistic scale. - Chain stage: the gated residual stream -> norm_out with the MLP stage's temb (end-to-end test now runs timestep -> ... -> norm_out); registry, report/figure hooks, doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
… and figure Produced by scripts/h3_evidence.py and scripts/h3_chain_replay.py --backward at 38d575f on a clean tree; rendered with scripts/plot_h3_evidence.py and embedded in the operator doc. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com> # Conflicts: # csrc/ops.cpp
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com> # Conflicts: # scripts/plot_h3_evidence.py
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Validate pinned tensor contents at preparation and load time. Execute chain prerequisites for selected stages and require checkpoint weights for weighted evidence. Reject invalid CUDA tile and backward inputs, align the CUDA wrapper guards, skip unsupported devices, and rotate comparative timing with raw samples. Validation: 143 H3 tests passed on B200 after a fresh CUDA build; illegal FP32 tiles fail compilation; isolated AdaLN replay and pinned-weight evidence generation passed; Black, isort, Ruff, and Flake8 passed. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Check timestep and modality bounds on the device before computing gather addresses, including unchecked calls and overflowing int64 indices. Reject non-prefix replay stage selections and partial backward chains before loading weights. Validation: rebuilt the CUDA extension on B200; tests/h3: 175 passed, including 12 isolated native bounds probes and 16 CLI regressions. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
CodeRabbit skips automatic reviews targeting non-default branches by default. Allow the test-h3 base branch so pushes to the H3 stack receive incremental reviews. Validation: configuration validated against the official CodeRabbit JSON schema. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Include the RFC 420 target branch in CodeRabbit automatic reviews so pushes to the H3 stack are reviewed instead of silently skipped. Validation: configuration parses and validates against the published CodeRabbit v2 schema; the additional branch pattern matches only test-h3. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
RFC RL-Align#420 row ws1_one_h3_block. Block 0 of the pinned checkpoint runs as a 17-node graph (AdaLN projection through the second gated residual) on packed FL2VA layouts that are bitwise equal to the pinned diffusers pipeline's. Each node runs as the diffusers replay, the RL-Kernel binding (chained and on provider inputs) and an FP32 golden, and names the RFC row that owns it and whether it is bound to that row's operator, an interim deterministic operator, or the provider replay. The report gives per node: repeat and batch-row bitwise checks, provider agreement, error against the golden for candidate and provider, sha256, and the first drift of each comparison; the backward replay gives leaf gradients (block input, temb, 12 parameters) and the first node whose incoming gradient drifts. A permutation probe reports control H-C3. The interim GEMM nodes use the Triton FP32-accumulate kernel with one BF16 rounding: TritonDetGemmOp's default tree path rounds every 32-wide leaf to BF16 and has about 3x cuBLAS's error at these shapes. Block-0 attention/FFN weights are pinned in a new manifest section and extracted with prepare_h3_weights.py --block. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Forward and backward replay of block 0 at S = 264, 925 and 3160, written from a clean tree at 41e4293: no repeat or batch-row drift at any node or gradient; first drift from diffusers is adaln_projection; block output as close to the FP32 golden as diffusers'. Adds the figure script. Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
|
Important Review skippedAuto reviews are disabled on base/target branches other than the default branch. Please check the settings in the CodeRabbit UI or the ⚙️ Run configuration
You can disable this status message by setting the Use the checkbox below for a quick retry:
Comment |
fusheng-ji
marked this pull request as ready for review
October 11, 2026 15:54
fusheng-ji
requested review from
Flink-ddd,
KJLdefeated,
bitborne,
inaniloquentee and
ryankert01
as code owners
October 11, 2026 15:54
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Status
No known failures at
41e4293(code) /caaf754(evidence). 11 of the 17 block nodes run on interim operators or the provider replay until their rows land (see Bindings). The block-level promises below hold with those bindings.What
RFC #420 row
ws1_one_h3_block: one real H3 block, forward and backward, with every intermediate and a first-drift report.MiniMax-H3@42ed227runs as a 17-node graph, from the AdaLN projection to the second gated residual. It runs on packed FL2VA layouts whoseposition_idsand tags are bitwise equal to the pinned diffusers pipeline'sbuild_packed_sequence; the hashes are pinned in a CPU test.first_{,isolated_,repeat_,batch_}drift.temband all 12 block parameters (repeat, error vs golden);first_grad_{repeat,batch}_drift.block_tensors). They are extracted byprepare_h3_weights.py --blockfrom the same shard, so the existing conditioning file is unchanged.Bindings
adaln_projectionadaln_projection_3modnorm1,norm2h3_rmsnormresidual_attn,residual_mlpadaln_gate_residualq/k/v_proj,o_proj,ffn_gate_up,ffn_downh3_qkv_gemm,h3_attention_o_gemm,h3_ffn_gate_up_gemm,h3_ffn_down_gemmq_norm,k_normh3_qk_rmsnorm_d128H3RMSNormCudaOpon the 128-wide heads (bitwise equal tonn.RMSNorm)rope_q,rope_kh3_mm_rope_3axis_partialattentionh3_full_attentionDeterministicAttentionOp(causal=False)swigluh3_swigluSwiGLUCudaOpWhen an owning row lands, swapping its node's
candidatecallable is the whole integration, and the report entries become that operator's block-level evidence.Prior art & reuse decision
This is a validation closeout, not a new kernel; no operator is written here.
MiniMaxH3TransformerBlockis the executable reference. It is replayed op for op as the provider (h3_provider), which pins its dtypes and op order without importing diffusers.rl_engine/validation/models/chain_gate.py). Its method is reused: per-node digests and first drift in graph order. Its graph is Qwen3's (causal, no AdaLN, Qwen shapes), so the H3 block is built on theh3_chainstage pattern of [WS1][CUDA][MiniMax-H3] timestep_sinusoid_h3: FP32 timestep features + shared H3 harness (RFC #420 stack 1/9) #483.TritonDetGemmOp.__call__reduces K as a tree with BF16 nodes; its FP32-accumulate kernel is equally deterministic and row-invariant, matches cuBLAS's accuracy, and is 9–17× faster (B200, error vs FP64):Results
Generated from a clean tree at
41e4293on an otherwise idle B200 (torch 2.13.0+cu130), withtools/validation/models/h3_block_replay.py --layouts tiny,small,medium:report.json.adaln_projection/adaln_projectionadaln_projection/adaln_projectionadaln_projection/adaln_projectionnorm1/norm2,q_norm/k_norm, RoPE, and both gated residuals. The first drift is the first reduction whose tree differs from cuBLAS.adaln_proj.weight,adaln_proj.biasandtemb. diffusers'index_selectbackward accumulates in BF16 atomics; the row operators use FP32 segment sums.attn.norm_q.weightat S = 3160 (9.5e-3 vs 5.5e-3). At S = 925 the same leaf is 2× more accurate than diffusers, so the 128-element gradient's error follows BF16 rounding, not one side.h3_full_attention.Tests
501 passed, 5 skipped on two B200s, including the 9 new tests in
test_h3_one_block.py. The new tests check:Scope / Known limitations
h3_mm_rope_3axis_partiallands.tembis the provider's (the conditioning chain is [WS1][CUDA][MiniMax-H3] timestep_sinusoid_h3: FP32 timestep features + shared H3 harness (RFC #420 stack 1/9) #483–[WS1][CUDA][MiniMax-H3] adaln_row_gather: fused six-way AdaLN row gather (RFC #420 stack 4/9) #486).ws1_full_h3_transformer.